Skip to main content

3. Saving and Loading Models

Saving a model means storing its learned weights so they can be used later.

Loading a model restores those learned weights.

Save Model​

The recommended approach is to save the model's state_dict.

torch.save(model.state_dict(), "model.pth")

Load Model​

Create the same model architecture first, then load the saved weights.

model = NeuralNetwork()
model.load_state_dict(torch.load("model.pth"))

For inference:

model.eval()

State Dictionary​

state_dict contains the model's learned parameters.

print(model.state_dict())

It mainly contains:

  • Weights
  • Biases

Save vs Load​

SaveLoad
torch.save()torch.load()
Stores model parametersRestores model parameters
Usually saves state_dictLoads into the same architecture

Model Checkpoint​

A checkpoint can store more than just model weights.

torch.save({
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"epoch": epoch,
"loss": loss
}, "checkpoint.pth")

Useful when training needs to be resumed later.

Save Only Weights vs Checkpoint​

Save state_dictSave checkpoint
Mainly for inference/reuseMainly for continuing training
Smaller and simpleStores training state too
Model architecture must be recreatedModel architecture and optimizer must be recreated

Actual Implementation​

import torch
import torch.nn as nn
import torch.optim as optim

X = torch.tensor([[1.0], [2.0], [3.0], [4.0]])
y = torch.tensor([[2.0], [4.0], [6.0], [8.0]])

class NeuralNetwork(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(1, 1)

def forward(self, x):
return self.linear(x)

model = NeuralNetwork()
optimizer = optim.SGD(model.parameters(), lr=0.01)
criterion = nn.MSELoss()

for epoch in range(1000):
output = model(X)
loss = criterion(output, y)

optimizer.zero_grad()
loss.backward()
optimizer.step()

# Save model
torch.save(model.state_dict(), "model.pth")

# Create a new model
loaded_model = NeuralNetwork()

# Load saved weights
loaded_model.load_state_dict(torch.load("model.pth", weights_only=True))

# Evaluation mode
loaded_model.eval()

with torch.no_grad():
predictions = loaded_model(X)

print(predictions)